Skip to main content

Chapter 6.1 - Multi-Head Self-Attention

Overview

Looking at a painting with just one eye is okay, but looking with both eyes gives you depth! Multi-Head Attention is like giving our AI lots of eyes, so it can look at a sentence from many different perspectives at the exact same time.


🎯 Why we do it

Rationale

Sometimes a word can mean different things. In "The bank of the river," bank means land. In "The bank was robbed," bank means money. If the AI only has one "Attention Head", it might get confused. By having multiple heads, one head can look for places, another can look for money, and another can look for action!

🛠️ How we do it

Methodology

Instead of doing the attention math once, we split our vectors into smaller pieces and give each piece to a different "Head". They all do their own little homework assignments independently. Once they are all done, we glue their answers back together!

# Imagine we have 8 dimensions
word = [1, 2, 3, 4, 5, 6, 7, 8]

# We split it for two heads!
head_1 = [1, 2, 3, 4] # Head 1 looks at the first half
head_2 = [5, 6, 7, 8] # Head 2 looks at the second half

# They do their math, and then we glue it back!
glued_together = head_1 + head_2

💻 Code Implementation

Here is the exact PyTorch implementation for the concepts discussed above:

import torch
import torch.nn as nn

class MultiHeadAttention(nn.Module):
def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_bias=False):
super().__init__()
assert d_out % num_heads == 0, "d_out must be divisible by num_heads"

self.d_out = d_out
self.num_heads = num_heads
self.head_dim = d_out // num_heads # Reduce the projection dim to match desired output dim

self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_key = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
self.out_proj = nn.Linear(d_out, d_out) # Linear layer to combine head outputs
self.dropout = nn.Dropout(dropout)
self.register_buffer("mask", torch.triu(torch.ones(context_length, context_length), diagonal=1))

def forward(self, x):
b, num_tokens, d_in = x.shape

keys = self.W_key(x) # Shape: (b, num_tokens, d_out)
queries = self.W_query(x)
values = self.W_value(x)

# We implicitly split the matrix by adding a `num_heads` dimension
# Unroll last dim: (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
keys = keys.view(b, num_tokens, self.num_heads, self.head_dim)
values = values.view(b, num_tokens, self.num_heads, self.head_dim)
queries = queries.view(b, num_tokens, self.num_heads, self.head_dim)

# Transpose: (b, num_tokens, num_heads, head_dim) -> (b, num_heads, num_tokens, head_dim)
keys = keys.transpose(1, 2)
queries = queries.transpose(1, 2)
values = values.transpose(1, 2)

# Compute scaled dot-product attention (aka self-attention) with a causal mask
attn_scores = queries @ keys.transpose(2, 3) # Dot product for each head

# Original mask truncated to the number of tokens and converted to boolean
mask_bool = self.mask.bool()[:num_tokens, :num_tokens]

# Use the mask to fill attention scores
attn_scores.masked_fill_(mask_bool, -torch.inf)

attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
attn_weights = self.dropout(attn_weights)

# Shape: (b, num_tokens, num_heads, head_dim)
context_vec = (attn_weights @ values).transpose(1, 2)

# Combine heads, where self.d_out = self.num_heads * self.head_dim
context_vec = context_vec.contiguous().view(b, num_tokens, self.d_out)
context_vec = self.out_proj(context_vec) # optional projection

return context_vec